Skip to content

[NNX] Delete Linen 5/5: remove the pure_nnx / enable_nnx / pure_nnx_decoder config flags - #4361

Merged
copybara-service[bot] merged 1 commit into
mainfrom
feat/nnx-del-linen-5-flags
Sep 11, 2026
Merged

copybara-service[bot] merged 1 commit into
mainfrom
feat/nnx-del-linen-5-flags

Conversation

@ecnal-cienet

Copy link
Copy Markdown
Collaborator

NNX Migration Route Map

  1. ✅ Add NNX scaffolding: pure_nnx flag, init_state_fn, TrainStateNNX, NNX utils. Linen workflow unchanged. (PR NNX migration prep (1/N): pure_nnx flag and init_state_fn scaffolding #3427)
  2. ✅ NNX sharding utilities: get_abstract_state_nnx, get_named_sharding_nnx, set_named_sharding_nnx, get_partition_spec_nnx, get_mesh_from_config. (PR NNX migration prep (2/N): NNX utils and sharding utilities #3470)
  3. ✅ NNX fully supported end-to-end: TrainStateNNX, model creation, gradient accumulation, checkpointing, and training loop dispatch. (PR NNX migration prep (3/N): TrainState, model creation, and end-to-end training loop #3500)
  4. ✅ Sharding diagnostics on NNX, plus post-training bugfixes that surfaced once the NNX path got exercised end-to-end. (PR [NNX] NNX migration prep (4/N): sharding tools and post-training fixes #3652)
  5. ✅ NNX correctness fixes, feature enablements, and vocab tiling on NNX.
  6. ✅ NNX-native DPO.
  7. ✅ NNX-native MaxEngine inference. (PR [NNX] NNX migration prep (7/N): NNX-native MaxEngine inference #3821)
  8. ✅ NNX-native LoRA + GRPO. (PR [NNX] NNX migration prep (8/N): NNX native lora grpo #3824)
  9. ✅ NNX-aware QK-Clip + remaining checkpoint utilities. (PR [NNX] NNX migration prep (9/N): NNX-aware QK-Clip + checkpoint utilities #3836)
    9.5. ✅ NNX + AQT in MaxEngine + serve-mode reload + gpt3 prefill fix. (PR [NNX] NNX migration prep (9.5/N): NNX + AQT in MaxEngine + serve-mode reload + gpt3 prefill fix #3844)
  10. ✅ Vocab tiling custom_vjp for NNX.
  11. ✅ Set NNX defaults to True; regenerate sharding goldens; flip back integration-test pure_nnx=False annotations. (PR [NNX] NNX migration (11/N): set pure_nnx / enable_nnx / pure_nnx_decoder defaults to True #3526)
  12. Delete Linen-specific code paths and NNX compatibility flags — delivered as 5 stacked PRs (12.1–12.5):
    12.1. ✅ Collapse Linen dispatch in utils / quantization / checkpointing / sharding. (PR [NNX] Delete Linen 1/5: collapse dispatch in utils, quantization, checkpointing, sharding #4357)
    12.2. ✅ Collapse Linen dispatch in trainers (pre-train, DiLoCo, GRPO); regenerate the NNX DPO golden. (PR [NNX] Delete Linen 2/5: collapse dispatch in trainers (pre-train, DiLoCo, GRPO) #4358)
    12.3. ✅ Collapse Linen dispatch in inference (maxengine, kvcache, vLLM, LoRA). (PR [NNX] Delete Linen 3/5: collapse dispatch in inference (maxengine, kvcache, vLLM, LoRA) #4359)
    12.4. ✅ Delete the Linen decoder/attention/embedding layers + *_as_linen model wrappers. (PR [NNX] Delete Linen 4/5: remove the Linen decoder/attention layers and *_as_linen model wrappers #4360)
    12.5. 🔄 [This PR] Remove the pure_nnx / enable_nnx / pure_nnx_decoder flags from types.py (Fields + dead qwix validator), base.yml, inference/vllm.yml, the distillation configs, and pyconfig_deprecated.py, plus doc/notebook/script mentions. No source or test reads the flags anymore; config init succeeds with them absent. Completes the Linen deletion (≈ −5.3k lines across the series); the original single-branch PR [NNX] NNX migration (12/N): delete Linen code paths, classes, and NNX compatibility flags #4038 can be closed. Stacks on 12.4.

Description

Final PR of the stacked "Delete Linen" series (see 1/5). With every read gone after 1–4, this removes the three migration flags — pure_nnx, enable_nnx, pure_nnx_decoder — from types.py (the Fields plus the now-dead qwix validator), base.yml, inference/vllm.yml, the distillation configs, and pyconfig_deprecated.py, and clears the remaining mentions in the docs, the LoRA demo notebook, and the distillation / gemma3-LoRA shell scripts.

Nothing reads the flags anymore (grep -rE 'pure_nnx|enable_nnx|pure_nnx_decoder' src/ tests/, excluding reference_hlo, is empty), and config init succeeds with them absent, including under use_qwix_quantization=True. After this the Linen deletion is complete — roughly −5.3k lines across the series — and the original single-branch PR #4038 can be closed. Base branch is feat/nnx-del-linen-4-layers. ~14 files, +9 / −54.

Checklist

Before submitting this PR, please make sure (put X in square brackets):

  • I have performed a self-review of my code. For an optional AI review, add the gemini-review label.
  • I have necessary comments in my code, particularly in hard-to-understand areas.
  • I have run end-to-end tests tests and provided workload links above if applicable.
  • I have made or will make corresponding changes to the doc if needed, including adding new documentation pages to the relevant Table of Contents (toctree directive) as explained in our documentation.

@codecov

codecov Bot commented Jul 6, 2026

Copy link
Copy Markdown

@ecnal-cienet
ecnal-cienet force-pushed the feat/nnx-del-linen-5-flags branch 2 times, most recently from d294837 to 75fb598 Compare July 7, 2026 21:00
@ecnal-cienet
ecnal-cienet force-pushed the feat/nnx-del-linen-5-flags branch 7 times, most recently from b5d2940 to 511d13a Compare July 21, 2026 19:11
@ecnal-cienet
ecnal-cienet force-pushed the feat/nnx-del-linen-5-flags branch 6 times, most recently from 1e6f41e to 83d2308 Compare July 29, 2026 14:48
@ecnal-cienet
ecnal-cienet force-pushed the feat/nnx-del-linen-5-flags branch 4 times, most recently from 8c1f026 to 7e9cd77 Compare August 20, 2026 21:38
@ecnal-cienet
ecnal-cienet force-pushed the feat/nnx-del-linen-5-flags branch 5 times, most recently from ce61748 to 6c3c9b6 Compare August 31, 2026 21:18
@ecnal-cienet
ecnal-cienet force-pushed the feat/nnx-del-linen-5-flags branch 2 times, most recently from 28a7c2b to 8dbdd97 Compare September 8, 2026 22:20
@ecnal-cienet
ecnal-cienet changed the base branch from main to feat/nnx-del-linen-4-layers September 8, 2026 23:07
@ecnal-cienet
ecnal-cienet force-pushed the feat/nnx-del-linen-4-layers branch from a4835ed to e93ede5 Compare September 8, 2026 23:10
@ecnal-cienet
ecnal-cienet force-pushed the feat/nnx-del-linen-5-flags branch from 8dbdd97 to e84b113 Compare September 8, 2026 23:10
…oder config flags

Drop the three migration flags from types.py (the Fields plus the now-dead qwix
validator), base.yml, inference/vllm.yml, the distillation configs and
pyconfig_deprecated.py, and clear the remaining mentions in the docs, the LoRA
demo notebook and the shell scripts.

setup_initial_state was the one reader parked by pre-train 1/3, because its
Linen branch was entangled with the checkpoint restore overlay. Collapse it
here. The Linen arm is unreachable now that get_abstract_state always returns an
nnx.State, and the weight_sparsity_n/m merge it carried is already covered by
_reshard_aligned, which keeps the freshly initialized value for any leaf the
checkpoint did not supply. The trailing unbox_logicallypartioned only unwrapped
Linen LogicallyPartitioned boxes, so it goes with it.

Removing pure_nnx_decoder also un-skips four train_compile pipeline tests that
were guarded on it. Three pass as they are, so drop all four guards.
test_pipeline_subset turned up a real porting bug: the NNX decoder tagged the
scan axis of the layers outside the pipeline "layers", which base.yml maps to
the stage mesh axis, so a remainder that does not divide the stage count failed
to shard. Linen used "layers_outside_pipeline", which has no rule; use that.

Tests follow: drop the flag from every config dict and SimpleNamespace stand-in
that still set it, and delete test_diloco_requires_pure_nnx, which asserted the
message of the validator this change removes.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants